代码实现 GQA Attention
从零理解 GQA Attention:结合 PyTorch 源码详解
首先先看下MiniMind的实现
def apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=1):
def rotate_half(x): return torch.cat((-x[..., x.shape[-1] // 2:], x[..., : x.shape[-1] // 2]), dim=-1)
q_embed = ((q * cos.unsqueeze(unsqueeze_dim)) + (rotate_half(q) * sin.unsqueeze(unsqueeze_dim))).to(q.dtype)
k_embed = ((k * cos.unsqueeze(unsqueeze_dim)) + (rotate_half(k) * sin.unsqueeze(unsqueeze_dim))).to(k.dtype)
return q_embed, k_embed
def repeat_kv(x: torch.Tensor, n_rep: int) -> torch.Tensor:
bs, slen, num_key_value_heads, head_dim = x.shape
if n_rep == 1: return x
return (x[:, :, :, None, :].expand(bs, slen, num_key_value_heads, n_rep, head_dim).reshape(bs, slen, num_key_value_heads * n_rep, head_dim))
class Attention(nn.Module):
def __init__(self, config: MiniMindConfig):
super().__init__()
self.num_key_value_heads = config.num_attention_heads if config.num_key_value_heads is None else config.num_key_value_heads
self.n_local_heads = config.num_attention_heads
self.n_local_kv_heads = self.num_key_value_heads
self.n_rep = self.n_local_heads // self.n_local_kv_heads
self.head_dim = config.head_dim
self.is_causal = True
self.q_proj = nn.Linear(config.hidden_size, config.num_attention_heads * self.head_dim, bias=False)
self.k_proj = nn.Linear(config.hidden_size, self.num_key_value_heads * self.head_dim, bias=False)
self.v_proj = nn.Linear(config.hidden_size, self.num_key_value_heads * self.head_dim, bias=False)
self.o_proj = nn.Linear(config.num_attention_heads * self.head_dim, config.hidden_size, bias=False)
self.q_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps)
self.k_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps)
self.attn_dropout = nn.Dropout(config.dropout)
self.resid_dropout = nn.Dropout(config.dropout)
self.dropout = config.dropout
self.flash = hasattr(torch.nn.functional, 'scaled_dot_product_attention') and config.flash_attn
def forward(self, x, position_embeddings, past_key_value=None, use_cache=False, attention_mask=None):
bsz, seq_len, _ = x.shape
xq, xk, xv = self.q_proj(x), self.k_proj(x), self.v_proj(x)
xq = xq.view(bsz, seq_len, self.n_local_heads, self.head_dim)
xk = xk.view(bsz, seq_len, self.n_local_kv_heads, self.head_dim)
xv = xv.view(bsz, seq_len, self.n_local_kv_heads, self.head_dim)
xq, xk = self.q_norm(xq), self.k_norm(xk)
cos, sin = position_embeddings
xq, xk = apply_rotary_pos_emb(xq, xk, cos, sin)
if past_key_value is not None:
xk = torch.cat([past_key_value[0], xk], dim=1)
xv = torch.cat([past_key_value[1], xv], dim=1)
past_kv = (xk, xv) if use_cache else None
xq, xk, xv = (xq.transpose(1, 2), repeat_kv(xk, self.n_rep).transpose(1, 2), repeat_kv(xv, self.n_rep).transpose(1, 2))
if self.flash and (seq_len > 1) and (not self.is_causal or past_key_value is None) and (attention_mask is None or torch.all(attention_mask == 1)):
output = F.scaled_dot_product_attention(xq, xk, xv, dropout_p=self.dropout if self.training else 0.0, is_causal=self.is_causal)
else:
scores = (xq @ xk.transpose(-2, -1)) / math.sqrt(self.head_dim)
if self.is_causal: scores[:, :, :, -seq_len:] += torch.full((seq_len, seq_len), float("-inf"), device=scores.device).triu(1)
if attention_mask is not None: scores += (1.0 - attention_mask.unsqueeze(1).unsqueeze(2)) * -1e9
output = self.attn_dropout(F.softmax(scores.float(), dim=-1).type_as(xq)) @ xv
output = output.transpose(1, 2).reshape(bsz, seq_len, -1)
output = self.resid_dropout(self.o_proj(output))
return output, past_kv
这段代码实际上就是一个 Llama/Qwen 风格的 GQA Attention 实现。
如果只看代码,很多人会迷失在:
view
transpose
repeat_kv
flash attention
cache
里面。
实际上整个流程只有 8 步:
输入 x
↓
Q投影
K投影
V投影
↓
分Head
↓
RoPE
↓
拼接KV Cache
↓
GQA扩展KV
↓
Attention
↓
输出投影
先从整体架构理解,再逐行拆解。
一、Attention初始化
配置参数
假设:
hidden_size = 512
num_attention_heads = 8
num_key_value_heads = 2
head_dim = 64
那么:
Q Head数 = 8
KV Head数 = 2
这就是 GQA。
代码:
self.n_local_heads = config.num_attention_heads
即:
8
self.n_local_kv_heads = self.num_key_value_heads
即:
2
计算:
self.n_rep = self.n_local_heads // self.n_local_kv_heads
得到:
8 // 2 = 4
表示:
每个KV Head
被4个Q Head共享
即:
Q0 Q1 Q2 Q3 -> KV0
Q4 Q5 Q6 Q7 -> KV1
这是 GQA 最核心的参数。
二、QKV Projection
输入:
x.shape = [B, L, hidden_size]
例如:
[2,128,512]
Query
xq = self.q_proj(x)
线性层:
512
↓
8 × 64
↓
512
输出:
[2,128,512]
Key
xk = self.k_proj(x)
注意:
num_key_value_heads=2
所以:
512
↓
2 × 64
↓
128
输出:
[2,128,128]
Value
同理:
xv = [2,128,128]
这就是 GQA 的第一处省显存。
传统 MHA:
Q=512
K=512
V=512
GQA:
Q=512
K=128
V=128
直接减少:
75%
三、拆成多个Head
Query
xq = xq.view(
bsz,
seq_len,
self.n_local_heads,
self.head_dim
)
变成:
[2,128,8,64]
即:
8个Q Head
Key
xk.view(...)
变成:
[2,128,2,64]
即:
2个KV Head
Value
同理:
[2,128,2,64]
此时:
Q : 8 heads
K : 2 heads
V : 2 heads
已经是标准 GQA 结构。
四、QK Norm
代码:
xq = self.q_norm(xq)
xk = self.k_norm(xk)
对应:
QK-Norm
这是 Llama3 引入的重要改进。
对每个 Head 做 RMSNorm:
然后:
作用:
避免QK数值爆炸
训练更稳定
五、RoPE
代码:
xq,xk=apply_rotary_pos_emb(...)
进入:
q_embed = q * cos + rotate_half(q) * sin
即:
其中:
本质:
把位置编码旋转进Q和K
最终:
Q包含位置信息
K包含位置信息
V不包含
六、KV Cache
推理时:
假设已经生成:
Hello world
缓存:
past_key_value =
(
old_k,
old_v
)
新token:
!
对应:
new_k
new_v
代码:
xk = torch.cat(
[past_key_value[0], xk],
dim=1
)
变成:
old_k + new_k
同理:
xv
拼接。
最终:
所有历史KV
保存在缓存中。
七、GQA核心
终于来到最关键部分。
当前:
xq = [B,L,8,64]
xk = [B,L,2,64]
xv = [B,L,2,64]
Head数对不上:
Q=8
K=2
Attention没法算。
于是:
repeat_kv(xk, n_rep=4)
进入:
x[:, :, :, None, :]
原来:
[2,128,2,64]
变:
[2,128,2,1,64]
expand:
expand(
bs,
slen,
2,
4,
64
)
变:
[2,128,2,4,64]
逻辑:
KV0
↓
复制4份
KV1
↓
复制4份
reshape:
[2,128,8,64]
变成:
KV0 KV0 KV0 KV0
KV1 KV1 KV1 KV1
对应:
Q0 → KV0
Q1 → KV0
Q2 → KV0
Q3 → KV0
Q4 → KV1
Q5 → KV1
Q6 → KV1
Q7 → KV1
这就是 GQA。
为什么这样不会增加KV Cache?
很多人第一次会疑惑:
你不是复制了吗?
实际上:
复制发生在:
当前Forward
阶段。
缓存中保存的仍然是:
[2,128,2,64]
而不是:
[2,128,8,64]
所以:
KV Cache大小仍然是2个Head
没有增加。
八、Transpose
代码:
xq.transpose(1,2)
从:
[B,L,H,D]
变:
[B,H,L,D]
得到:
Q = [2,8,128,64]
Attention实现都喜欢:
Batch
Head
Seq
Dim
这种布局。
九、Attention计算
Flash Attention路径
F.scaled_dot_product_attention(...)
直接调用:
PyTorch FlashAttention
内部融合:
QKᵀ
Softmax
AV
速度最快。
十、普通Attention路径
Step1
scores = (xq @ xk.transpose(-2,-1)) / sqrt(head_dim)
得到:
[B,H,L,L]
即:
Attention Score
数学上:
Step2
因果Mask
scores += upper_triangle(-inf)
变成:
未来token不可见
例如:
I love
不能偷看:
you
Step3
Softmax
A = softmax(scores)
对应:
Step4
乘V
output = A @ xv
对应:
输出:
[B,H,L,D]
十一、合并Head
代码:
output.transpose(1,2)
回到:
[B,L,H,D]
然后:
reshape(
bsz,
seq_len,
-1
)
变:
[B,L,8×64]
即:
[2,128,512]
十二、输出投影
output = self.o_proj(output)
对应:
即 Transformer 最后一层:
Concat Heads
↓
Linear
↓
Hidden Size
输出:
[2,128,512]
整个GQA流程图

从源码实现角度看,GQA 真正新增的代码其实只有 repeat_kv() 这一处。其它流程(QKV 投影、RoPE、Mask、Softmax、输出投影)与普通 MHA 基本一致。
因此可以把 GQA 理解为:
训练和推理时只存储少量 KV Head,在计算 Attention 的瞬间,通过
repeat_kv将 KV 逻辑上扩展到与 Query Head 数量一致,从而兼顾 KV Cache 节省和多头表达能力。
这也是为什么 Llama 2/3、Qwen、DeepSeek 等模型都采用 num_heads > num_kv_heads 的设计。真正节省显存的不是 Attention 计算本身,而是长期保存的 KV Cache。